FlashAttention: GPU 메모리 계층을 이해한 최초의 IO 인식 정확 어텐션 알고리즘
2026. 9. 2.
연구 배경 및 핵심 문제 의식
트랜스포머의 셀프 어텐션은 시퀀스 길이 N에 대해 시간과 메모리 모두 O(N²)의 복잡도를 가진다. 시퀀스가 길어질수록 이 이차적 증가는 빠르게 실용적 한계에 봉착하며, 특히 긴 문서, 고해상도 이미지, 장기 시계열 데이터를 다루는 태스크에서 심각한 병목이 된다.
이 문제를 해결하기 위해 수많은 근사 어텐션 방법들이 제안되었다. 린포머, 리포머, 롱포머, 빅버드, 퍼포머 등 다양한 접근들은 FLOP 수를 선형 또는 준선형으로 줄이는 데 집중했다. 그러나 이들 대부분은 이론적 FLOP 절감에도 불구하고 실제 wall-clock 속도 향상을 보이지 못했다.
"Many approximate attention methods have aimed to reduce the compute and memory requirements of attention. These methods range from sparse-approximation [51, 74] to low-rank approximation [84, 50, 12], and their combinations [3, 92, 9]. Although these methods reduce the compute requirements to linear or near-linear in sequence length, many of them do not display wall-clock speedup against standard attention and have not gained wide adoption. One main reason is that they focus on FLOP reduction (which may not correlate with wall-clock speed) and tend to ignore overheads from memory access (IO)."
— 원문 링크
FlashAttention의 저자들은 이 역설의 원인을 GPU 메모리 계층 구조에 대한 무지에서 찾았다. 현대 GPU에서 연산 처리량(FLOP/s)은 빠르게 성장했지만 HBM의 메모리 대역폭은 그에 비해 훨씬 느리게 성장했다. 따라서 어텐션처럼 메모리 바운드 특성이 강한 연산에서는 FLOP 감소보다 HBM 접근 횟수 감소가 실제 속도에 훨씬 결정적이다. FlashAttention은 이 원리를 명시적으로 알고리즘 설계에 적용한 최초의 IO 인식 정확 어텐션이다.
선행 연구와의 패러다임 비교
표준 어텐션은 Q, K, V를 HBM에서 읽어 N×N 크기의 어텐션 행렬 S와 P를 HBM에 기록한 뒤 다시 읽어 최종 출력을 계산한다. 이 과정에서 Θ(Nd + N²)의 HBM 접근이 발생하며, N이 커질수록 N² 항이 지배적이 된다.
기존 근사 어텐션 방법들은 N² 항을 줄이기 위해 희소 근사나 저순위 근사를 활용했다. 그러나 이들 역시 중간 행렬의 HBM 입출력 오버헤드를 충분히 줄이지 못했고, 짧은 시퀀스 구간에서는 오히려 표준 어텐션보다 느린 경우도 있었다.
FlashAttention은 정확한 어텐션 값을 유지하면서도 HBM 접근 횟수를 Θ(N²d²M⁻¹)으로 낮추며, 이는 실제 d와 M 값에서 표준 어텐션 대비 최대 9배 적은 HBM 접근에 해당한다.
"Let N be the sequence length, d be the head dimension, and M be size of SRAM with d≤M≤Nd. Standard attention (Algorithm 0) requires Θ(Nd+N2) HBM accesses, while FlashAttention (Algorithm 1) requires Θ(N2d2M−1) HBM accesses."
— 원문 링크
핵심 아키텍처 및 구현 실증
FlashAttention의 핵심 아이디어는 두 가지 고전적 기법의 결합이다. 하나는 타일링을 통해 전체 입력을 보지 않고도 점진적으로 소프트맥스를 계산하는 것이고, 다른 하나는 역방향 패스에서 중간 어텐션 행렬을 저장하는 대신 온칩에서 재연산하는 것이다. 이 두 기법이 시너지를 내어 HBM 접근을 극적으로 줄이면서도 수치적으로 동일한 정확 어텐션을 계산할 수 있게 한다.
타일링과 온라인 소프트맥스: HBM 접근의 핵심 절감 기법
소프트맥스는 전체 행의 최댓값과 합산을 알아야 정규화할 수 있어 본질적으로 전역 연산처럼 보인다. FlashAttention은 온라인 소프트맥스 기법을 활용해 이 문제를 해결한다. 입력 Q, K, V를 SRAM 크기에 맞게 블록으로 분할하고, 각 블록을 순회할 때마다 현재까지의 최댓값 m과 합산 통계량 ℓ을 갱신하면서 출력 O를 점진적으로 보정한다.
이 타일링 전략 덕분에 N×N 크기의 어텐션 행렬 S와 P를 HBM에 기록하거나 다시 읽을 필요가 없다. 모든 중간 계산이 온칩 SRAM 안에서 완결되며, HBM과의 통신은 최소화된다. HBM 접근 횟수가 어텐션 런타임의 주된 결정 요인임은 실험적으로도 검증되었다.
"We validate that the number of HBM accesses is the main determining factor of attention run-time. In Fig. 2 (left), we see that even though FlashAttention has higher FLOP count compared to standard attention (due to recomputation in the backward pass), it has much fewer HBM accesses, resulting in much faster runtime. In Fig. 2 (middle), we vary the block size Bc of FlashAttention, which results in different amounts of HBM accesses, and measure the runtime of the forward pass. As block size increases, the number of HBM accesses decreases (as we make fewer passes over the input), and runtime decreases. For large enough block size (beyond 256), the runtime is then bottlenecked by other factors (e.g., arithmetic operations). Moreover, larger block size will not fit into the small SRAM size."
— 원문 링크
FlashAttention은 추가적으로 CUDA 커널 하나에 어텐션 관련 모든 연산(행렬 곱, 소프트맥스, 드롭아웃, 마스킹)을 커널 퓨전하여 구현한다. 드롭아웃과 마스킹이 함께 사용될 때 이 커널 퓨전 효과로 인해 A100 GPU에서 일반적으로 2~4배의 속도 향상이 관측된다.
"Figure 5 shows speedup on an A100 GPU with batch size 8, head dimension 64, and 12 attention heads, across different sequence lengths. We generally see 2-4× speedup, and we see more speedup when using dropout and masking due to kernel fusion."
— 원문 링크
역방향 패스 재연산: O(N) 메모리로 정확한 그래디언트 계산
표준 어텐션의 역방향 패스는 그래디언트를 계산하기 위해 순전파에서 생성된 N×N 크기의 어텐션 행렬 S와 P를 HBM에 저장했다가 다시 읽어야 한다. 이것이 표준 어텐션이 O(N²) 메모리를 요구하는 주된 이유다.
FlashAttention은 순전파 시 소프트맥스 정규화 통계량 (m, ℓ)만 HBM에 저장한다. 이는 O(N) 공간으로 충분하다. 역방향 패스에서는 이 통계량을 활용해 Q, K, V 블록으로부터 S와 P를 온칩 SRAM에서 빠르게 재연산한다. 재연산으로 인해 FLOPs는 더 많아지지만, HBM 접근 횟수가 줄어들기 때문에 오히려 역방향 패스가 더 빨라진다.
"Recomputation. One of our goals is to not store O(N2) intermediate values for the backward pass. The backward pass typically requires the matrices 𝐒,𝐏∈ℝN×N to compute the gradients with respect to 𝐐,𝐊,𝐕. However, by storing the output 𝐎 and the softmax normalization statistics (m,ℓ), we can recompute the attention matrix 𝐒 and 𝐏 easily in the backward pass from blocks of 𝐐,𝐊,𝐕 in SRAM. This can be seen as a form of selective gradient checkpointing [34, 10]. While gradient checkpointing has been suggested to reduce the maximum amount of memory required [66], all implementations (that we know off) have to trade speed for memory. In contrast, even with more FLOPs, our recomputation speeds up the backward pass due to reduced HBM accesses (Fig. 2). The full backward pass description is in Appendix B."
— 원문 링크
FlashAttention은 입출력 이외의 추가 메모리를 시퀀스 길이에 대해 선형인 O(N)만 요구한다. 이는 기존 그래디언트 체크포인팅 방법들이 속도를 희생해서 메모리를 줄이던 것과 달리, 속도와 메모리 모두를 동시에 개선한 핵심 차별점이다.
"Algorithm 1 returns 𝐎=softmax(𝐐𝐊⊤)𝐕 with O(N2d) FLOPs and requires O(N) additional memory beyond inputs and output."
— 원문 링크
IO 복잡도 최적성 증명과 블록 스파스 확장
논문은 FlashAttention의 HBM 접근 횟수가 단순히 개선된 수준을 넘어 이론적 하한임을 증명한다. 정확한 어텐션을 계산하는 어떤 알고리즘도 모든 SRAM 크기 M에 대해 FlashAttention의 HBM 접근 횟수인 o(N²d²M⁻¹)보다 점근적으로 개선할 수 없다.
"Let N be the sequence length, d be the head dimension, and M be size of SRAM with d≤M≤Nd. There does not exist an algorithm to compute exact attention with o(N2d2M−1) HBM accesses for all M in the range [d,Nd]."
— 원문 링크
나아가 블록 희소 플래시어텐션으로 확장하면 블록 희소성 마스크의 비영 블록 비율 s에 비례하여 IO 복잡도가 추가로 개선된다. LRA 벤치마크에서 블록 스파스 FlashAttention은 표준 어텐션과 유사한 정확도를 유지하면서 2.8배의 속도 향상을 달성한다.
블록 스파스 FlashAttention은 기존에 알려진 모든 근사 어텐션 방법보다 모든 시퀀스 길이에서 빠른 런타임을 달성한다.
"Figure 3 (left) reports the runtime in milliseconds of the forward + backward pass of FlashAttention and block-sparse FlashAttention compared to the baselines in exact, approximate, and sparse attention (exact numbers in Appendix E). Runtime grows quadratically with sequence length, but FlashAttention runs significantly faster than exact attention baselines, up to 3× faster than the PyTorch implementation. The runtimes of many approximate/sparse attention mechanisms grow linearly with sequence length, but FlashAttention still runs faster than approximate and sparse attention for short sequences due to fewer memory accesses. The approximate attention runtimes begin to cross over with FlashAttention at sequences between 512 and 1024. On the other hand, block-sparse FlashAttention is faster than all implementations of exact, sparse, and approximate attention that we know of, across all sequence lengths."
— 원문 링크
벤치마크 성능: 속도, 메모리, 그리고 모델 품질
BERT-large(시퀀스 길이 512) 학습에서 FlashAttention은 MLPerf 1.1 학습 속도 기록을 보유했던 NVIDIA 구현보다 15% 빠른 단일 노드 최고 학습 속도를 달성했다.
"FlashAttention yields the fastest single-node BERT training speed that we know of. We train a BERT-large [22] model with FlashAttention on Wikipedia. Table 1 compares our training time to the implementation from Nvidia that set the training speed record for MLPerf 1.1 [58]. Our implementation is 15% faster."
— 원문 링크
GPT-2 학습에서는 HuggingFace 구현 대비 최대 3배, Megatron-LM 대비 최대 1.7배의 종단 간 학습 속도 향상이 관측되었다. 수치적 안정성 측면에서도 기존 구현과 동일한 검증 손실 곡선을 보임이 확인되었다.
"FlashAttention yields faster training times for GPT-2 [67] on the large OpenWebtext dataset [32] than the widely used HuggingFace [87] and Megatron-LM [77] implementations. Table 2 shows up to 3× end-to-end speedup compared to Huggingface and 1.7× speedup compared to Megatron-LM. FlashAttention achieves the same perplexity as the other two implementations, as we do not change the model definition. Appendix E includes plots of the validation perplexity throughout training, confirming that FlashAttention is as numerically stable as the baselines and produces the same training / validation curves."
— 원문 링크
FlashAttention이 가져오는 가장 인상적인 결과는 모델 품질의 향상이다. 컨텍스트 길이를 4K로 늘린 GPT-2는 Megatron이 컨텍스트 길이 1K로 학습하는 것보다 30% 더 빠르면서도 perplexity가 0.7 낮다. 긴 문서 분류 태스크에서는 6.4 포인트의 성능 향상이 관측된다.
"The runtime and memory-efficiency of FlashAttention allow us to increase the context length of GPT-2 by 4× while still running faster than the optimized implementation from Megatron-LM. Table 4 shows that that GPT-2 with FlashAttention and context length 4K is still 30% faster than GPT-2 from Megatron with context length 1K, while achieving 0.7 better perplexity."
— 원문 링크
가장 극적인 결과는 Path-X 태스크에서 나온다. 시퀀스 길이 16K에서 FlashAttention은 기존 트랜스포머 중 최초로 랜덤 이상의 성능인 61.4% 정확도를 달성했다. 블록 희소 플래시어텐션은 시퀀스 길이를 64K까지 확장해 Path-256에서 63.1% 정확도를 기록하며 기존 어떤 시퀀스 모델도 해결하지 못했던 벤치마크를 처음으로 돌파했다.
"The Path-X and Path-256 benchmarks are challenging tasks from the long-range arena benchmark designed to test long context. The task is to classify whether two points in a black and white 128×128 (or 256×256) image have a path connecting them, and the images are fed to the transformer one pixel at a time. In prior work, all transformer models have either run out of memory, or only achieved random performance [80]. There has been a search for alternative architectures that can model such long context [37]. We present here the first result of Transformer models being able to solve Path-X and Path-256 (Table 6). We pretrain a transformer on Path-64, and then transfer to Path-X by spatially interpolating the positional embeddings. FlashAttention achieves 61.4 accuracy on Path-X. Additionally, block-sparse FlashAttention enables the Transformers to scale to sequence length 64K, achieving 63.1 accuracy on Path-256."
— 원문 링크
메모리 측면에서도 FlashAttention의 메모리 사용량은 시퀀스 길이에 대해 선형으로 증가하며, 표준 정확 어텐션 대비 최대 20배 메모리 효율적이다. 64K 길이에서도 린포머보다 2배 더 효율적이다.
"Figure 3 (right) shows the memory footprint of FlashAttention and block-sparse FlashAttention compared to various exact, approximate, and sparse attention baselines. FlashAttention and block-sparse FlashAttention have the same memory footprint, which grows linearly with sequence length. FlashAttention is up to 20× more memory efficient than exact attention baselines, and is more memory-efficient than the approximate attention baselines. All other algorithms except for Linformer run out of memory on an A100 GPU before 64K, and FlashAttention is still 2× more efficient than Linformer."
— 원문 링크
실무적 한계 및 후속 연구 팩트체크
FlashAttention은 단일 GPU의 HBM 접근 최적화에서 이론적 하한을 달성했지만, 이 최적성은 단일 디바이스 환경을 전제로 한다. 수백~수천 GPU를 활용하는 분산 학습 환경에서는 GPU 간 통신 병목이 새로운 지배적 요인으로 등장하며, HBM 최적화만으로는 이 문제를 해결할 수 없다.
역방향 패스 재연산 전략도 한계가 있다. FlashAttention v1은 스레드 블록 간 작업 분배와 GPU 워프 스케줄링을 충분히 최적화하지 못해, 이론적으로 달성 가능한 최대 GPU 활용률에 미치지 못했다. 이는 FlashAttention-2가 해결한 핵심 과제이기도 하다.
블록 희소 어텐션 확장은 GPU 아키텍처와 모델 구조에 대한 강한 가정을 전제로 한다. 다양한 하드웨어 환경이나 비정형 어텐션 패턴에서는 블록 희소성 마스크의 하드웨어 친화적 실행이 보장되지 않는다.